自定义环境
ROLL 的 Agentic Pipeline 使用 GEM 环境。一个环境是注册到某个 env_type 的 gem.Env 实现;环境管理器通过 gem.make() 创建实例,并不断把模型输出传给 step()。
本文描述 roll/pipeline/agentic/env 与 roll/pipeline/agentic/env_manager 当前采用的契约。旧版文档中的 BaseDiscreteActionEnv / BaseLanguageBasedEnv 分类已经不再使用。
运行流程
标准 TrajEnvManager 中的单个 Episode 流程如下:
gem.make(env_type, **env_config)
-> reset(seed)
-> 根据 observation 和 info["env_instruction"] 构造提示词
-> 模型生成动作字符串
-> step(action_string)
-> 重复,直到 terminated 或达到 max_steps
环境实例运行在 Environment Worker 中,并且可能并发执行。不要依赖可变的全局状态,并保证 reset(seed) 可复现。若环境持有外部资源,应实现可重复调用的 close();具体是在每个 Episode 后还是退出时调用取决于所选环境管理器,因此有严格生命周期要求的环境还应在终止和异常路径中清理资源。
环境接口契约
继承 gem.Env 并实现以下方法:
from typing import Any
from gem import Env
class MyEnv(Env):
def __init__(self, max_steps: int = 10, **kwargs): ...
def reset(self, seed: int | None = None) -> tuple[Any, dict[str, Any]]: ...
def step(
self, action: str
) -> tuple[Any, float, bool, bool, dict[str, Any]]: ...
def close(self) -> None: ...
__init__(**env_config)
custom_envs.<tag>.env_config 下的配置会由下面的代码直接传给构造函数:
gem.make(env_id=env_config["env_type"], **env_config["config"])
构造函数应只接收有意义的选项并尽早校验。包装第三方环境时,可以保留 **kwargs 以传递其参数。
reset(seed)
适用时先调用 GEM 基类实现,重置 Episode 的所有状态,然后返回:
(observation, info)
observation是当前要展示给 Agent 的状态;对TrajEnvManager而言通常是字符串。info必须是字典。可将稳定的任务说明或动作格式要求放入info["env_instruction"],标准管理器会在第一轮将其加入提示词。- 仅在没有可用 Episode 时返回
(None, info);管理器会将其视为本次没有 Rollout。 - 所有影响任务生成的随机源都应使用传入的 seed。同一个 Rollout Group 内的环境会收到相同的 Episode seed。
step(action)
标准管理器传入的是模型解码后的完整响应,而不是离散动作 ID。环境需自行解析并校验文本,然后返回 Gymnasium 风格的五元组:
(observation, reward, terminated, truncated, info)
observation:动作执行后的状态。reward:本轮的标量奖励;ROLL 会将各轮奖励相加得到 Episode score。terminated:任务因成功或失败等自然终止条件结束。truncated:因超时等外部限制结束。在当前标准TrajEnvManager中,只有truncated=True不会停止循环;环境自身发生超时时,应同时返回terminated=True, truncated=True。达到custom_envs.<tag>.max_steps时,管理器也会强制结束并设置 truncation。info:本轮附加信息。
内置环境通常使用以下可选 info 字段:
info = {
"action_desc": "动作结果的可读描述",
"metrics": {
"action_is_valid": True,
"action_is_effective": True,
"success": False,
},
"metrics_agg_mode": {
"action_is_valid": "mean",
"action_is_effective": "mean",
"success": "last",
},
# "suffix": "当 agent_template 包含 {suffix} 时渲染的额外状态",
}
metrics 的值应为数值或布尔值,metrics_agg_mode 指定 ROLL 如何在整条轨迹上聚合各项指标。遇到格式错误的动作时,通常保持 observation 不变、返回格式惩罚、设置 action_is_valid=False,并允许 Episode 继续。
Observation 格式
TrajEnvManager 期望文本类 observation,并由管理器构造对话历史。AgentNativeStepEnvManager 用于由环境自行维护对话的场景,此类环境返回 OpenAI 风格的消息列表,例如 [{'role': 'user', 'content': '...'}]。环境和管理器必须成对选择,不要混用两种契约。使用 Tool Call 的环境还应配置对应的 ToolCallRunner 或 Native Runner。
多模态 Observation
环境需要返回图片或视频时,使用 roll.pipeline.agentic.env_manager.vl_traj_env_manager.VLTrajEnvManager。它沿用相同的 reset() 和五元组 step() 契约,但支持以下 observation 形式:
str:纯文本的一轮。numpy.ndarray:单张 RGB 图片。数组必须能传给PIL.Image.fromarray(obs, mode="RGB"),通常为H x W x 3的uint8数组。dict:多模态的一轮。prompt保存文本或 Chat 风格内容,image和/或video保存模型 Transformers processor 能接收的媒体对象。
典型的图片 observation 如下:
from PIL import Image
class VisualQuestionEnv(Env):
image_placeholder = "<image>"
def reset(self, seed=None):
super().reset(seed)
image = Image.open(self.image_path).convert("RGB")
observation = {
"prompt": "<image>\nWhat object is highlighted?",
"image": [image],
}
return observation, {
"env_instruction": "Answer directly, or request a visual tool action."
}
字典的 key 使用单数形式:image 和 video。值可以是单个对象或 list/tuple;建议始终使用列表,当 prompt 中包含多个媒体占位符时也必须使用列表。媒体值的顺序必须与 prompt 中占位符的顺序一致。
如果 prompt 使用显式占位符,应声明对应的类属性:
class MyMultimodalEnv(Env):
image_placeholder = "<image>"
video_placeholder = "<video>"
VLTrajEnvManager 会将环境占位符替换成 DataCollatorWithPaddingForMM 要求的特殊 token。它按照轨迹顺序累积所有轮次的媒体,使用策略模型的 ProcessorMixin 构造模型输入,并将处理后的多模态字段保留在 Rollout 中。环境应该返回原始媒体和文本,不要自行返回预计算的 pixel_values、token ID 或设备上的 tensor。
需要结构化 Chat 内容时,prompt 也可以是 content item 列表:
observation = {
"prompt": [
{"type": "text", "text": "Inspect this image:"},
{"type": "image"},
{"type": "text", "text": "What changed?"},
],
"image": [image],
}
视觉语言环境需要配置 VLTrajEnvManager 及其两个轮次模板:
custom_envs:
VisualQuestion:
env_type: visual_question
env_manager_cls: roll.pipeline.agentic.env_manager.vl_traj_env_manager.VLTrajEnvManager
max_steps: 4
max_tokens_per_step: 256
agent_system_template: "You are a visual reasoning agent."
pre_step_template: "\nTurn {turn_idx}:\n"
next_step_template: |
You have {actions_left} actions left.
Keep the response within {max_response_length} tokens.
env_config:
image_path: /path/to/image.png
Actor 必须是拥有兼容 tokenizer/processor 的视觉语言模型。VLTrajEnvManager 当前只为图片和视频提供了明确的 Collation 路径;音频或新的媒体类型需要扩展管理器和多模态 Collator。媒体加载可能失败时,应在 reset()/step() 中校验并返回受控的终止结果,避免 prompt 与错误媒体静默配对。
最小实现
import random
import re
from typing import Any
from gem import Env
class GuessNumberEnv(Env):
def __init__(self, low: int = 1, high: int = 10, format_penalty: float = -0.1, **kwargs):
self.low = low
self.high = high
self.format_penalty = format_penalty
self.target = None
self.done = False
def reset(self, seed: int | None = None) -> tuple[str, dict[str, Any]]:
super().reset(seed)
self.target = random.Random(seed).randint(self.low, self.high)
self.done = False
return (
f"Guess an integer from {self.low} to {self.high}.",
{"env_instruction": "Reply with <answer>number</answer>."},
)
def step(self, action: str) -> tuple[str, float, bool, bool, dict[str, Any]]:
match = re.search(r"<answer>\s*(-?\d+)\s*</answer>", action)
valid = match is not None
guess = int(match.group(1)) if valid else None
success = valid and guess == self.target
self.done = success
if not valid:
observation, reward = "Invalid format; try again.", self.format_penalty
elif guess < self.target:
observation, reward = "Too small.", 0.0
elif guess > self.target:
observation, reward = "Too large.", 0.0
else:
observation, reward = "Correct.", 1.0
info = {
"metrics": {"action_is_valid": valid, "success": success},
"metrics_agg_mode": {"action_is_valid": "mean", "success": "last"},
}
return observation, reward, self.done, False, info
def close(self) -> None:
pass
注册环境
在环境包中导出该类,并在 roll/pipeline/agentic/env/__init__.py 中添加 GEM 延迟注册:
gem.register(
"guess_number",
entry_point="roll.pipeline.agentic.env.guess_number:GuessNumberEnv",
)
注册名必须与 env_type 完全一致。不要在包的 __init__.py 中导入重量级或可选依赖;只有 gem.make() 创建环境时才应加载 entry point。若环境依赖可选组件,可参考现有可选环境的受保护注册方式。
配置环境
在 custom_envs 下添加一个 tag(examples/config 中部分公共配置片段使用单数 custom_env,由上层配置合并到 custom_envs):
env_manager_cls: roll.pipeline.agentic.env_manager.traj_env_manager.TrajEnvManager
custom_envs:
GuessNumber:
env_type: guess_number
env_manager_cls: ${env_manager_cls}
agent_runner_cls: null
max_steps: 8
max_tokens_per_step: 32
agent_system_template: "You are a careful game-playing agent."
agent_template: |
Turn {turn_idx}:
Observation: {observation}
You have {actions_left} actions left.
Respond with one action only.
env_config:
low: 1
high: 20
format_penalty: -0.1
train_env_manager:
num_env_groups: 32
group_size: 4
tags: [GuessNumber]
num_groups_partition: [32]
配置分为两层:
max_steps、max_tokens_per_step、模板、env_manager_cls和agent_runner_cls等字段由 Rollout 框架使用。- 只有
env_config内的字段会传给环境构造函数。 tags按custom_envs的 key 选择配置,env_type则选择已注册的 GEM 类。num_groups_partition必须与tags一一对应,且总和等于num_env_groups;同一 Group 的成员共享配置和 seed。
初始化或执行步骤不是线程安全时,设置 use_thread_lock: true。昂贵的共享后端可用 max_env_step_concurrent 限制并发调用数。
验证清单
在完整训练前先验证环境契约:
import roll.pipeline.agentic.env # 执行 ROLL 的 GEM 注册
import gem
env = gem.make(env_id="guess_number", low=1, high=3)
obs, info = env.reset(seed=42)
assert isinstance(info, dict)
obs, reward, terminated, truncated, info = env.step("<answer>2</answer>")
assert isinstance(reward, (int, float))
assert isinstance(terminated, bool) and isinstance(truncated, bool)
assert isinstance(info, dict)
env.close()
还需检查:
- 相同 seed 是否生成相同初始任务;
- 合法动作、非法动作、自然终止、超时和最大步数路径;
- 每次
step()是否严格返回五个值; - 各步的 metric key 和聚合模式是否保持一致;
- 异常和
close()后是否正确释放资源; - 使用所选 Environment Manager 与 Runner 的小规模 Rollout 是否能完整运行。